AdamWeightDecay

对权重执行 Adam Weight Decay(AdamW)优化更新,原地写回 varmv

\[\begin{split}\begin{aligned} m_t &= m_{t-1} + (g_t - m_{t-1}) \cdot (1 - \beta_1) \\ v_t &= v_{t-1} + (g_t^{2} - v_{t-1}) \cdot (1 - \beta_2) \\ var_t &= var_{t-1} - lr \cdot \left( \frac{m_t}{\sqrt{v_t} + \epsilon} + decay \cdot var_{t-1} \right) \end{aligned}\end{split}\]

其中 \(g_t\) 为梯度,\(m_t\) / \(v_t\) 为一阶 / 二阶动量, \(decay\) 为权重衰减系数。本实现对区间 [start, end) 内元素逐点更新。

输入:
  • var - 待更新权重张量地址(输入/输出,原地更新)

  • m - 一阶动量张量地址(输入/输出,原地更新)

  • v - 二阶动量张量地址(输入/输出,原地更新)

  • grad - 梯度张量地址

  • float_Parameters - 浮点超参数组,长度 5,布局见下

  • int_Parameters - 整数参数数组,长度 2,布局见下

  • core_mask - 核掩码(仅共享存储版本使用)

float_Parameters 布局:

  • [0] lr - 学习率

  • [1] beta1 - 一阶动量衰减系数

  • [2] beta2 - 二阶动量衰减系数

  • [3] epsilon - 数值稳定项

  • [4] decay - 权重衰减系数

int_Parameters 布局:

  • [0] start - 参与计算的起始索引(含)

  • [1] end - 参与计算的结束索引(不含)

输出:
  • var / m / v - 原地写回更新结果

支持平台:

FT78NE MT7004

备注

  • FT78NE 支持 fp32

  • MT7004 支持 fp16、fp32

  • hp_ 版本中 float_Parameters 元素类型为 float16fp_ 版本为 float

共享存储版本:

void hp_adam_weight_decay_s(float16 *var, float16 *m, float16 *v, float16 *grad, float16 *float_Parameters, int *int_Parameters, int core_mask)
void fp_adam_weight_decay_s(float *var, float *m, float *v, float *grad, float *float_Parameters, int *int_Parameters, int core_mask)

C调用示例:

 1// MT7004 示例(共享存储多核,DDR 地址)
 2void TestAdamWeightDecaySMCFp32(int length, int core_mask) {
 3    int core_id = get_core_id();
 4    int logic_core_id = GetLogicCoreId(core_mask, core_id);
 5    int core_num = GetCoreNum(core_mask);
 6    float *var = (float *)0x81000000;
 7    float *m = (float *)0x82000000;
 8    float *v = (float *)0x83000000;
 9    float *grad = (float *)0x84000000;
10    float *float_Parameters = (float *)0x85000000;
11    int *int_Parameters = (int *)0x85000014;
12    if (logic_core_id == 0) {
13        float_Parameters[0] = 0.001f;  // lr
14        float_Parameters[1] = 0.9f;    // beta1
15        float_Parameters[2] = 0.999f;  // beta2
16        float_Parameters[3] = 1e-4f;   // epsilon
17        float_Parameters[4] = 0.01f;   // decay
18        int_Parameters[0] = 0;         // start
19        int_Parameters[1] = length;   // end
20    }
21    sys_bar(0, core_num);
22    fp_adam_weight_decay_s(var, m, v, grad, float_Parameters, int_Parameters, core_mask);
23}
24
25void main() {
26    int core_mask = 0b1111;
27    TestAdamWeightDecaySMCFp32(4096, core_mask);
28}

私有存储版本:

void hp_adam_weight_decay_p(float16 *var, float16 *m, float16 *v, float16 *grad, float16 *float_Parameters, int *int_Parameters)
void fp_adam_weight_decay_p(float *var, float *m, float *v, float *grad, float *float_Parameters, int *int_Parameters)

C调用示例:

 1// MT7004 示例(私有存储单核,AM 地址)
 2void TestAdamWeightDecayAMFp32(int length) {
 3    float *var = (float *)0x10000000;
 4    float *m = (float *)0x10020000;
 5    float *v = (float *)0x10040000;
 6    float *grad = (float *)0x10060000;
 7    float *float_Parameters = (float *)0x1007F000;
 8    int *int_Parameters = (int *)0x1007F014;
 9    float_Parameters[0] = 0.001f;  // lr
10    float_Parameters[1] = 0.9f;    // beta1
11    float_Parameters[2] = 0.999f;  // beta2
12    float_Parameters[3] = 1e-4f;   // epsilon
13    float_Parameters[4] = 0.01f;   // decay
14    int_Parameters[0] = 0;         // start
15    int_Parameters[1] = length;   // end
16    fp_adam_weight_decay_p(var, m, v, grad, float_Parameters, int_Parameters);
17}
18
19void main() {
20    TestAdamWeightDecayAMFp32(2048);
21}